Skip to content

feat(attention): fixed-m splash attention kernel with dynamic bounds and safety fallbacks - #477

Open
Perseus14 wants to merge 1 commit into
mainfrom
feat/fixed-m-kernel
Open

Perseus14 wants to merge 1 commit into
mainfrom
feat/fixed-m-kernel

Conversation

@Perseus14

@Perseus14 Perseus14 commented Sep 13, 2026 •

Copy link
Copy Markdown
Collaborator

Summary

Fixed-m splash attention kernel. When a provable bound on every logit (‖q‖·‖k‖, optionally centred on the mean key k̄) fits the fp32 exponent range, the kernel uses that fixed max m instead of an online running max. This removes the per-block max/rescale work from the softmax inner loop. When the bound can't be proven, it falls back to the standard online-softmax kernel (_lse_scan).

What's in it

  • custom_splash_attention.py: fixed-m kernel with dynamic Q/KV bounds. The in-kernel q·k̄ dot runs at Precision.HIGHEST, and _splash_attention_forward_ring rejects 2D (2, num_q_heads) mk metadata in favor of 3D (2, num_q_heads, num_q_blocks).
  • Contracts:
    • fixed_m_norms_squared is a required argument.
    • value=None gives v_ok=False (fail-closed).
    • The V-max check runs before the all-to-all; only Q and K go through the optimization barrier.
    • _custom_ring_attention_forward validates fixed_m_norms[0] shape against per_q_block.
  • Guards: short KV (< flash_min_seq_length) falls back to dot_product, cross-attention remaps to tokamax_flash, and heads % num_shards != 0 raises ValueError.

Known limitations

  • On the uncentred R > 1 ring path the fixed bound can still flush the numerator of heads far below it. This is not measured.
  • The ring kernel runs its own pmin over global-bound eligibility unless the caller passes uniform_fixed_m=True.

Tests

custom_splash_fixed_m_test.py (31) and ring_fixed_m_test.py (18):

  • kernel relative-L2 checks, mk shape and V-metadata fail-closed checks;
  • fixed_m_norms_squared and v_ok contract tests (ring_fixed_m_test.py);
  • uniform and hybrid branches, d=64 gate check, ragged KV tails (N%8==0, N%bkv≠0, odd-multiple-of-8 tail), and a U=2/R=2 pre-all-to-all V-outlier check;
  • 6 ring numerical tests vs an f32 reference + 1 squared-vs-unsquared parity test + 3 raw-vs-centered bound tests.
  • CI vs Local VM: 18 of the 49 tests (kernel numerics grids and multi-device integration) skip when GITHUB_ACTIONS=true to keep CI fast; end_to_end/tpu/run_wan_stack_tests.sh runs all 49 on a TPU VM.
  • Results (TPU v6e-8): 49 passed (run_wan_stack_tests.sh, 170.4 s).

Stack: #477 → #478 → #479 → #488 → #491

@github-actions

Copy link
Copy Markdown

@gemini-code-assist gemini-code-assist Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Code Review

This pull request introduces dynamic fixed-m constants and safe bounds calculation based on KV sequence length, adds support for virtual K-centering, and enhances input validation and dtype safety checks in the Pallas flash attention kernel. It also expands the test suite to cover various edge cases, including per-Q-block fallback, batched isolation, and virtual K-centering. The review feedback focuses on optimizing the TPU kernel performance by replacing expensive dynamic integer division with optimized BlockSpec mapping and introducing a static use_k_centering flag to conditionally compile the centering logic at trace time.

Comment thread src/maxdiffusion/kernels/custom_splash_attention.py
Comment thread src/maxdiffusion/kernels/custom_splash_attention.py
Comment thread src/maxdiffusion/kernels/custom_splash_attention.py Outdated
Comment thread src/maxdiffusion/kernels/custom_splash_attention.py Outdated
@Perseus14
Perseus14 force-pushed the feat/fixed-m-kernel branch 10 times, most recently from 3b9d7aa to cb1e4d1 Compare September 15, 2026 06:37
@syhuang22
syhuang22 self-requested a review September 15, 2026 20:04

@syhuang22 syhuang22 left a comment •

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

A few things before this can go in though:

  • If this lands on its own, fixed-m gets turned off for R>1 (use_fixed_m = False in _ulysses_ring_custom_attention, it only comes back in #478). That breaks the recipe we ship today. Can we split the stack so each PR is safe to merge by itself?
  • Cutting the Ulysses gate from 213 to 113 feels more conservative than we need. With centered keys, the mass you can lose is bounded by 2^(ceil(U)-C-126), so a gate around C+116 already keeps the loss under 0.1%. Could you share fallback rates so we can pick the number?
  • mk[0] means something different now, but a 2D mk still gets silently broadcast. Can we just raise on that?
  • The V check never fired in my runs on WAN 2.2 (every call passed), and it only exists because C moved from 88 to 102. Is it worth the extra complexity?
  • Please put back the comments that explain the Mosaic cliff and k-smoothing. They save the next person a lot of pain.

@syhuang22 syhuang22 left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Some line-level notes to go with my comment above.

Comment thread src/maxdiffusion/models/attention_flax.py Outdated
Comment thread src/maxdiffusion/kernels/custom_splash_attention.py
Comment thread src/maxdiffusion/kernels/custom_splash_attention.py Outdated
Comment thread src/maxdiffusion/kernels/custom_splash_attention.py
Comment thread src/maxdiffusion/kernels/custom_splash_attention.py
Comment thread src/maxdiffusion/models/attention_flax.py
Comment thread src/maxdiffusion/models/attention_flax.py Outdated
Comment thread src/maxdiffusion/models/attention_flax.py
Comment thread src/maxdiffusion/kernels/splash_attention/ring_attention_kernel.py
Comment thread src/maxdiffusion/generate_ltx_video.py
@Perseus14
Perseus14 force-pushed the feat/fixed-m-kernel branch 4 times, most recently from 44dfe63 to fc91941 Compare September 17, 2026 12:31
syhuang22
syhuang22 previously approved these changes Sep 17, 2026
@Perseus14
Perseus14 force-pushed the feat/fixed-m-kernel branch 3 times, most recently from d57775a to 678e806 Compare September 17, 2026 18:57
@Perseus14
Perseus14 requested a review from eltsai September 17, 2026 19:09
@Perseus14 Perseus14 self-assigned this Sep 17, 2026
@Perseus14
Perseus14 added this pull request to stack #486 September 17, 2026 19:10
Comment thread src/maxdiffusion/models/attention_flax.py
Comment thread src/maxdiffusion/kernels/splash_attention/ring_attention_kernel.py Outdated
Comment thread src/maxdiffusion/models/attention_flax.py
…and safety fallbacks

Implements exact fixed-m splash attention in Pallas on TPU. Instead of an
online running max, a head / Q-block whose Cauchy-Schwarz bound
max||q|| * max||k - k_bar|| is within a length-dependent window uses a
precomputed m = ceil(bound) - C(N) (+ q . k_bar with virtual K-centering) and
skips the per-block rescales; everything else keeps online softmax.

- `get_fixed_m_constants(N)`: length-dependent recenter C(N) and safe bound
  that keep the f32 numerator/denominator from overflowing or flushing.
- `fixed_m_dtype_is_safe`: rejects fp16/fp8 (exponent range too small for the
  recentered exponent); bf16 and fp32 pass.
- V-magnitude gate: |V| <= DEFAULT_MAX_V_BOUND (256), otherwise online.
- Per-(head, Q-block) dispatch in the kernel; the Ulysses and R=1 ring callers
  compile a uniform (all-fixed) and a hybrid kernel and pick one with
  `lax.cond` on `all_fixed`. Both kernels are compiled; the extra compile
  time / program size has not been measured.
- Virtual K-centering: the kernel receives k_mean and adds q . k_bar to the
  fixed bound in VMEM (an MXU dot_general on an (Hkv, 8, d) tiled k_mean,
  precision=HIGHEST because the gate has no slack at its edge). K itself is
  never re-written in HBM. k_mean is padded to the lane-padded head_dim after
  `_pad_data_for_flash`, so head_dim < 128 works (regression test at d=64).

Explicit metadata contracts on the fixed-m ring path
----------------------------------------------------
1. Norm representation is declared, not inferred. The gate previously guessed
   whether `fixed_m_norms` were squared with `(qn.max() * mk.max()) < 1000.0`.
   Magnitude cannot answer that: legacy unsquared norms of (1000, 2) have a
   true bound of 2000, but read as squared give sqrt(2000) ~= 44.7, a ~45x
   under-estimate that admits fixed-m where it must fall back. Replaced by a
   REQUIRED `fixed_m_norms_squared` argument (no default; omitting it with
   fixed-m on raises). Tests use squared norms.

2. The V-safety predicate is required, not assumed. The V-magnitude and dtype
   verdict cannot be re-derived from one hop's Q/K, so the ring kernel raises
   unless `v_ok` is passed (`value=None` in the metadata helper gives
   v_ok = 0, fail-closed). `_ulysses_ring_custom_attention` (R > 1) computes it
   on the PRE-all-to-all shards -- dtype safety plus |V| <= 256 -- and reduces
   it with pmin over BOTH internal axes (ring, ulysses): before the a2a each
   device holds a sequence shard of every head, and the ring's
   accumulate-vs-LSE branch must be taken uniformly by every ppermute
   participant. Only Q and K go through the optimization_barrier that keeps
   XLA from duplicating their producer chains into the norm reductions; V is
   reduced outside it so the Q/K a2a does not wait on V's producer.

3. Norm shape is validated against per_q_block. Both gates compute
   `qn * mk[:, None]`, so a (num_heads,) array supplied with per_q_block=True
   broadcasts to (num_heads, num_heads) and pairs head j's query norm with head
   h's key norm; a sink head then inherits a small bound and the kernel
   evaluates exp2(large_logit - small_m) -> inf. The kernel now rejects the
   mismatch.

Known limitations
- On the uncentered R > 1 ring path the fixed bound can still flush the
  numerator for heads far below the bound; this is not measured here.
- The ring kernel runs its own pmin over global-bound eligibility (added by
  this commit) unless the caller promises `uniform_fixed_m=True`; on the
  Ulysses+Ring path that pmin is redundant with the caller's (not measured).

Tests
- custom_splash_fixed_m_test.py (31): kernel unit tests with relative-L2
  checks, k-norm (mk) shape and V-metadata fail-closed checks, integration
  through `_ulysses_attention` and the R=1 ring (uniform branch, hybrid
  branch via a sink head, d=64 -- the test re-derives the gate from the
  inputs and asserts it admits fixed-m, it does not observe the cond -- ragged
  KV tails with N % 8 == 0 and N % block_kv != 0 whose tail is an odd multiple
  of 8), and a U=2/R=2 check that a single V outlier on one pre-a2a device
  makes v_ok False on every device (ring kernel stubbed; needs >= 4 devices,
  explicit skip otherwise).
- ring_fixed_m_test.py (18): 8 backend-independent contract tests (including
  the required fixed_m_norms_squared and v_ok arguments), 6 ring numerical
  tests vs an f32 reference (2 devices), 1 test that declared-unsquared and
  squared norms give the same result, and 3 tests showing the raw-K bound is
  unsound where the centered bound is not.
- CI: to keep GitHub Actions short, 18 of these 49 tests (the kernel
  numerics grids and the multi-device integration tests) skip when
  GITHUB_ACTIONS=true. The 9 tests from main and the contract, dtype and
  V-metadata tests still run there. The new
  end_to_end/tpu/run_wan_stack_tests.sh runs all of them on a TPU VM; later
  PRs in this stack add their test files to it.

Verified (final tree): TPU v6e-8 (jax 0.11.2, end_to_end/tpu/run_wan_stack_tests.sh): 49 passed, 6 subtests passed in 170.4s.
@Perseus14
Perseus14 force-pushed the feat/fixed-m-kernel branch from 53905e1 to 4ec211c Compare October 2, 2026 17:15
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants